前幾天的文章介紹如何做 vector search,今天介紹如何在 vector search 的基礎上加入 Scalar (metadata) filtering,並聊一些背後的設計思考。
Milvus 中的 Filtered Search 分為Standard filtering 標準篩選和Iterative filtering 迭代篩選
Standard filtering 的好處是可以把搜尋範圍縮小到一個很小的範圍,但對於複雜的篩選條件來說 (ex: Json 條件、字串比較),Scalar Filtering 的本身可能才是效能瓶頸。
Iterative filtering 則改為逐步從 ANN 搜尋結果中取得候選,再對候選執行 Scalar Filtering。通常表示更少的 Scalar filtering,但當篩選條件過於嚴格時,反而會導致更長的處理時間或潛在的效能問題。
程式範例參考:
# Standard filtering
milvus_client.search(
collection_name="my_collection",
data=[query_vector],
limit=5,
filter='color like "red%" and likes > 50',
output_fields=["color", "likes"]
)
# Iterative filtering
milvus_client.search(
collection_name="my_collection",
data=[query_vector],
limit=5,
filter='color like "red%" and likes > 50',
output_fields=["color", "likes"],
search_params={
"hints": "iterative_filter"
}
)
如果對前幾天提到的 ANN 還有印象的話,這時候應該會有一個疑問,篩選過資料如何在既有的 Index 上去做搜尋呢?尤其是 graph-based index,我們不可能用篩選後的資料去建新的 HNSW 呀。
在 Standard filtering 中,篩選會建立一個 bitset(位元集合),用來標記哪些向量符合過濾條件,如果在搜尋時把不符合條件的都直接排除掉,會破壞 graph 的連通性,可能導致 recall 大幅下降;反過來說,如果無論 metadata 是否符合都訪問所有節點,那麼搜尋空間就會急遽膨脹,效能也會大幅降低。
在 Milvus 的 blog 整理了一些他們在 filtered graph search 的研究方向,雖然沒有明確說明具體實作方式,但也幫助我們一窺其背後的一些設計思考:
假設原本的 graph 長這樣:A → B → C → D
假設 B 或 C 在篩選後被排除,那連結就斷了,A 永遠走不到 D,形成 island 孤島。
所以不能粗暴地直接把不符合條件的節點都拿掉,而是要想辦法決定「哪些被 filter 掉的節點,仍值得暫時走過」。
α 允許被經過當作搜尋路徑,其中 α 依實際的 filtering ratio 決定。不只按照「向量相似度」建圖,也替常用的 metadata 條件另外建立更適合過濾搜尋的圖。
簡單來說,除了一個 Base Graph 外,也替 metadata filter 建立 Column Graphs (Sub Graphs)。
但問題來了,難道每個 metadata value 都建一張完整 graph?如果每個值都建完整 graph 成本會爆掉,所以會限制每個 point 的 out-degree,讓 graph 變小。
最後一個可能是最值得研究的,也是最多應用程式的現況,是 scalar (metadata) 本來就有一個自己的關聯式資料庫,同時又有一個向量資料庫,一般流程可能是:
但資料量大(條件多)時可能會遇到效能瓶頸,所以 Milvus 做了一層 overhead (Milvus Server 需支援 Search Iterator V2 API),可以用 search_iterator() 搭配 external_filter_func 來實現類似 Iterative filtering 的效果:
from pymilvus.client.search_result import Hits
from pymilvus import MilvusClient
import numpy as np
valid_ids = {1, 12, 123, 1234} # external search result
def external_filter_func(hits: Hits):
return [
hit for hit in hits
if hit.id in valid_ids
]
milvus_client = MilvusClient(uri="http://localhost:19530")
milvus_client.load_collection("documents")
dim = 384 # adjust according to vector dimension
query_vector = np.random.rand(dim).astype(np.float32).tolist()
iterator = milvus_client.search_iterator(
collection_name="documents",
data=[query_vector],
batch_size=10,
anns_field="vector",
output_fields=["text"],
external_filter_func=external_filter_func,
)
results = []
try:
while True:
hits = iterator.next()
if not hits:
break
results.extend(hits)
if len(results) >= 10:
results = results[:10]
break
finally:
iterator.close()
凡事往好處想,想不到就算了。